🌳 EAGLE‑2 动态草稿树 并行解码 · 置信度驱动 · 动态剪枝

一种在 EAGLE-2 中提出的自适应草稿生成机制,通过 上下文感知的树结构,在保证生成质量的同时最大化投机解码的加速比。

🧠 核心概念 · 一句话说清

动态草稿树 = 每步根据 置信度 动态决定 “扩展哪个分支、扩展多深”,把算力优先花在最有希望的候选 token 上,然后 剪枝 后提交给目标模型并行验证。

🎯 要解决的问题

静态草稿树(EAGLE‑1)固定每层分支数,但不同上下文的接受率差异巨大——简单位置可以多猜几个,困难位置猜多了全是废的。静态结构无法适应,浪费算力。

💡 核心思想

  • 上下文感知:用草稿模型的 置信度分数(softmax 概率)近似目标模型会给出的接受率。
  • 动态扩展:每层只保留 Top‑K 个高置信度 token 继续扩展,低置信度的直接放弃。
  • 硬约束:总节点数不超过 speculative_num_draft_tokens(默认 8),保证验证开销可控。

🔧 简化网络结构 + 推理过程示例

以下用 极简版 EAGLE‑2 网络数值示例 展示动态草稿树的完整推理过程。

📐 极简网络结构

+----------------------------------------------------------------------+ | EAGLE‑2 草稿模型(简化) | +----------------------------------------------------------------------+ | | | 输入: hidden_states (目标模型上一层的输出) | | + 当前 token_ids | | ↓ | | [ 轻量级 Transformer 层 × 1 ] | | ↓ | | 输出: logits (词汇表大小的概率分布) | | + confidence (每个 token 的置信度分数) | | | | 关键组件: | | ┌─────────────────────────────────────────────────────────────┐ | | │ • 特征投影层: hidden_size → 中间维度 (如 1024) │ | | │ • 自注意力层: 多头注意力 (4头) │ | | │ • 输出层: 中间维度 → vocab_size │ | | │ • 置信度头: 中间维度 → 1 (sigmoid 输出 0~1 置信度) │ | | └─────────────────────────────────────────────────────────────┘ | | | +----------------------------------------------------------------------+

⚡ 实际 EAGLE‑2 草稿模型通常为 2‑4 层 Transformer,远小于目标模型(如 7B/13B),因此生成开销极低。

📝 推理过程示例(生成 “The cat” 后续)

📌 场景设置

  • 当前上下文:“The cat” (已生成 2 个 token)
  • 目标:生成后续 3‑4 个 token(如 “ sat on the mat”)
  • 参数max_nodes=8top_k=3max_depth=4
🔹 步骤 1:根节点扩展(生成第 1 层)

草稿模型基于 “The cat” 的 hidden_states,并行预测前 3 个候选 token:

候选 token置信度路径概率
sat0.850.85
slept0.600.60
jumped0.400.40

✅ 选择 Top‑3(top_k=3)全部保留,进入下一层。

🔹 步骤 2:扩展第 2 层(动态分支)

对每个候选 token,继续预测下一个 token:

父节点候选子节点置信度路径概率 (累积)
sat on0.900.85 × 0.90 = 0.765
sat at0.500.85 × 0.50 = 0.425
slept peacefully0.700.60 × 0.70 = 0.420
slept on0.450.60 × 0.45 = 0.270
jumped over0.550.40 × 0.55 = 0.220
jumped up0.350.40 × 0.35 = 0.140

✅ 保留路径概率最高的 max_nodes=8 个节点(此处共有 6 个,全部保留,继续扩展)。

🔹 步骤 3:扩展第 3 层 & 剪枝

继续扩展,直到节点数达到 max_nodes=8,然后剪枝:

路径(累积概率排序)累积概率状态
sat → on → the0.765 × 0.80 = 0.612✅ 保留 (Top-1)
sat → on → a0.765 × 0.30 = 0.230✅ 保留
slept → peacefully → on0.420 × 0.75 = 0.315✅ 保留
sat → at → the0.425 × 0.65 = 0.276✅ 保留
slept → on → the0.270 × 0.70 = 0.189✅ 保留
jumped → over → the0.220 × 0.60 = 0.132✅ 保留
jumped → up → the0.140 × 0.40 = 0.056❌ 剪枝 (概率最低)
sat → on → mat0.765 × 0.10 = 0.077❌ 剪枝 (概率低)

✅ 最终保留 6 条路径(≤8),提交给目标模型验证。

🔹 步骤 4:目标模型验证

将 6 条路径拼接成 batch,并行送入目标模型(如 LLaMA‑2 7B):

路径验证结果推进的 token
sat → on → the✅ 全部接受 sat, on, the
slept → peacefully → on❌ 第1个 slept 即被拒
sat → at → thesat 接受,❌ at 被拒 sat
slept → on → the❌ 第1个 slept 被拒
jumped → over → the❌ 第1个 jumped 被拒
sat → on → asat 接受,❌ on 被拒 sat

✅ 本轮实际推进:3 个 token sat + on + the),仅用了 1 次大模型前向。

🎯 步骤 5:总结
  • 总候选节点:从 1 个根节点动态扩展到 6 个叶子节点(≤8)。
  • 大模型前向次数:仅 1 次(验证整棵树)。
  • 实际推进3 个 token(接受率 50%)。
  • 对比自回归:如果让大模型自己生成 3 个 token,需要 3 次前向 → EAGLE‑2 加速约 3 倍

❓ 三个你最关心的问题

1️⃣ 每次推理到底蹦多少个 token(总节点数)?

不是固定数量,而是树的总节点数固定,形状动态。 总节点数由 speculative_num_draft_tokens 控制,默认值为 8(部分模型默认 4)。

公式关系: speculative_num_draft_tokens = speculative_num_steps × topk + 1 (近似,实际受剪枝影响)
例:topk=4num_steps=3 → 未剪枝前 1 + 4 + 16 + 64 = 85 个候选,但强制剪枝到 ≤8 个节点 再提交验证。

结论:EAGLE‑2 每次提交给目标模型验证的 草稿节点总数固定(如 8 个),但树的深度和每层宽度根据上下文实时变化

2️⃣ 输入和输出是什么?

角色输入输出
整体当前已生成的 token 序列(prompt + 已生成部分)被目标模型接受的 token(可能 1~8 个),追加到序列中
草稿模型hidden_states(目标模型上一轮输出的隐藏状态)每个候选 token 的 置信度分数 + logits
目标模型草稿树中的所有 token 序列(拼接成 batch)每个位置的 接受/拒绝 决策,并顺带输出新 logits

3️⃣ 动态草稿树 vs 静态草稿树

对比维度静态草稿树 (EAGLE‑1)动态草稿树 (EAGLE‑2)
树结构固定深度、固定每层分支数动态,根据置信度每层自适应
分支选择全量扩展,无筛选每层只保留 Top‑K 高置信度候选
剪枝策略无(或简单按位置截断)综合概率剪枝,保留路径概率最高的 N 条
适应能力弱,困难样本浪费算力,简单样本深搜,困难样本宽搜
加速比~2‑3x~3‑5x(比 EAGLE‑1 提升 20‑40%)

📊 完整决策流程图

graph TD A["开始本轮草稿生成
输入: hidden_states"] --> B["并行预测所有候选 token
(一次前向输出 K 个候选及其置信度)"] B --> C{"对当前层每个节点
计算上下文置信度分数"} C --> D["选择 Top‑K 个高置信度 token
作为下一层的候选分支"] D --> E["将选中的分支加入树结构
(记录路径概率)"] E --> F{"是否达到最大深度
或总节点数 ≥ max_nodes?"} F -->|否| B F -->|是| G["剪枝:按路径累积概率排序
只保留 Top‑N 条完整路径"] G --> H["将剪枝后的树(≤max_nodes 个节点)
拼接成 batch 提交给目标模型"] H --> I["目标模型并行验证所有节点
输出每个位置的接受/拒绝"] I --> J["接受有效 token,推进生成序列"] J --> K["更新 hidden_states 用于下一轮"] K --> L["结束本轮验证"] style A fill:#e3f0ff,stroke:#4a8fe7 style D fill:#fef3c7,stroke:#d4a24e style G fill:#fde2e2,stroke:#c9635a style H fill:#d1fae5,stroke:#3f9a75 style J fill:#d1fae5,stroke:#3f9a75 style L fill:#eef2f6,stroke:#9aa9bb
⚙️ 关键步骤:置信度驱动 动态分支扩展 剪枝降开销 并行验证

💻 源码参考(来自 EAGLE 官方实现)

以下为 EAGLE‑2 动态草稿树构建的核心 Python 逻辑(基于 SafeAILab/EAGLE 简化示意)。

# ============================================================
# EAGLE‑2 动态草稿树构建(核心逻辑)
# ============================================================

def eagle2_generate_tree(
    draft_model,
    hidden_states,          # 目标模型上一轮输出的隐藏状态
    token_ids,              # 当前已生成的 token 序列
    max_nodes = 8,       # 总节点数上限(提交给目标模型验证的草稿数量)
    top_k = 4,            # 每层保留的候选分支数
    max_depth = 4          # 最大树深度
):
    # 1. 初始化:根节点(当前已生成的位置)
    tree = [{
        'tokens': token_ids,
        'hidden': hidden_states,
        'prob': 1.0,           # 路径累积概率
        'depth': 0
    }]

    # 2. 循环扩展,直到达到最大节点数或最大深度
    while len(tree) < max_nodes and tree[-1]['depth'] < max_depth:
        current_nodes = [node for node in tree if node['depth'] == tree[-1]['depth']]
        if not current_nodes:
            break

        # 3. 并行前向:一次 forward 得到所有候选 token 的 logits 和置信度
        logits, confidence = draft_model.forward(current_nodes)

        # 4. 对每个节点,选择 Top‑K 个最高置信度的 token 作为下一层分支
        next_nodes = []
        for idx, node in enumerate(current_nodes):
            top_indices = confidence[idx].topk(top_k).indices.tolist()
            for token_idx in top_indices:
                # 计算路径累积概率(用于后续剪枝)
                path_prob = node['prob'] * confidence[idx][token_idx].item()
                next_nodes.append({
                    'tokens': node['tokens'] + [token_idx],
                    'hidden': logits[idx][token_idx],  # 简化:实际为 hidden 更新
                    'prob': path_prob,
                    'depth': node['depth'] + 1
                })

        # 5. 剪枝:如果总节点数超过 max_nodes,只保留累积概率最高的 max_nodes 条路径
        tree += next_nodes
        if len(tree) > max_nodes:
            tree.sort(key=lambda x: x['prob'], reverse=True)
            tree = tree[:max_nodes]

    # 6. 返回树结构(每个节点是一条候选路径)
    return tree


# ============================================================
# 目标模型验证阶段(简略)
# ============================================================
def verify_with_target(target_model, tree):
    # 将所有候选路径拼接成 batch,并行送入目标模型
    batch = [node['tokens'] for node in tree]
    outputs = target_model.forward(batch)
    # 逐位置判断接受/拒绝(树验证 + 接受率计算)
    accepted_tokens = []
    for i, tokens in enumerate(batch):
        if outputs[i].accepted:   # 该路径被接受
            accepted_tokens.append(tokens[-1])  # 取最后一个 token
    return accepted_tokens
            
⚠️ 上述代码为原理性示意,实际实现包含 KV Cache 管理、树结构验证、多路径回溯等复杂逻辑。完整源码请参考 EAGLE 官方仓库

⚙️ 关键参数与默认值

参数默认值作用
speculative_num_draft_tokens8目标模型一次验证的 总节点数(草稿数量上限),影响延迟和接受率
top_k4每层扩展时保留的 候选分支数,越大树越宽
max_depth4树的最大深度,限制最长依赖路径
confidence_threshold0.3置信度截断阈值,低于此值的 token 直接丢弃(可选)

💡 调参直觉top_k 大 → 树更宽,覆盖更多候选,但总节点数固定(max_nodes=8),所以会压缩深度max_depth 大 → 树更深,能生成更长的草稿,但浅层剪枝更激进。

📈 动态树 vs 静态树 · 工作流对比

📦 静态树 (EAGLE‑1)

  • 固定每层 top_k,所有分支均匀扩展
  • 不考虑上下文难度,简单/困难样本一视同仁
  • 复杂样本 → 大量低质量草稿,验证浪费
  • 简单样本 → 受限于固定深度,无法充分挖掘潜力
⏱️ 加速比:~2‑3x

🚀 动态树 (EAGLE‑2)

  • 每层根据 置信度 动态选择 Top‑K
  • 简单位置 → 增加深度,多猜几步
  • 困难位置 → 增加宽度,多试几种可能
  • 总节点数固定(max_nodes=8),算力利用最大化
⚡ 加速比:~3‑5x(比静态提升 20‑40%)
EAGLE‑2 动态草稿树 · 完整技术解析 · 基于 arXiv:2404.13249 与官方源码